# -*- coding: utf-8 -*-
"""
Q2: Spatial Differentiation Pattern and Native-Origin Effect Analysis
Based on precise annotation of university location information
Revision: save_all_data function adds df_raw parameter for full traceability

Required package installation:
    pip install pandas openpyxl matplotlib numpy scipy openpyxl
"""
import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
import matplotlib
import re
from collections import defaultdict
from scipy import stats
import os
import warnings
warnings.filterwarnings('ignore')

# SCI-style global plotting settings: Times New Roman serif font
matplotlib.rcParams['font.family'] = 'serif'
matplotlib.rcParams['font.serif'] = ['Times New Roman', 'DejaVu Serif']
matplotlib.rcParams['font.size'] = 11
matplotlib.rcParams['axes.labelsize'] = 12
matplotlib.rcParams['axes.titlesize'] = 14
matplotlib.rcParams['legend.fontsize'] = 10
matplotlib.rcParams['figure.titlesize'] = 14
matplotlib.rcParams['xtick.labelsize'] = 10
matplotlib.rcParams['ytick.labelsize'] = 9
matplotlib.rcParams['axes.spines.top'] = False
matplotlib.rcParams['axes.spines.right'] = False

# ==================== 1. Data Loading ====================
def load_data(file_path):
    df = pd.read_excel(file_path, sheet_name='汇总')
    df = df.iloc[:, [1, 2, 3, 4]].copy()
    df.columns = ['项目序号', '项目名称', '申报地区', '设课大学']
    df = df.dropna(subset=['项目名称'])
    return df

# ==================== 2. Core Function: Precisely Parse Province from Annotation ====================
def parse_province_from_annotation(school_entry):
    """
    Precisely extract province from the "University Name (Province City)" format
    Example: "Qilu Normal University (Shandong Jinan)" -> "山东省"
             "Beijing Sport University (Beijing Beijing)" -> "北京市"
    """
    if pd.isna(school_entry) or not school_entry:
        return ''
    match = re.search(r'[（(]([^）)]+)[）)]', str(school_entry))
    if not match:
        return ''
    location_str = match.group(1).strip()
    if location_str.startswith('内蒙古'):
        return '内蒙古自治区'
    elif location_str.startswith('广西'):
        return '广西壮族自治区'
    elif location_str.startswith('西藏'):
        return '西藏自治区'
    elif location_str.startswith('宁夏'):
        return '宁夏回族自治区'
    elif location_str.startswith('新疆'):
        return '新疆维吾尔自治区'
    elif len(location_str) >= 2:
        prov = location_str[:2]
        if prov in ['北京', '上海', '天津', '重庆']:
            return prov + '市'
        else:
            return prov + '省'
    return ''

def parse_school_and_province(school_entry):
    if pd.isna(school_entry) or not school_entry:
        return '', ''
    school_entry = str(school_entry).strip()
    match = re.match(r'^([^（(]+)', school_entry)
    school_name = match.group(1).strip() if match else school_entry
    province = parse_province_from_annotation(school_entry)
    return school_name, province

def parse_schools_from_row(schools_str):
    if pd.isna(schools_str) or str(schools_str).strip() == '无':
        return []
    parts = re.split('[、，,;；]', str(schools_str))
    results = []
    for part in parts:
        part = part.strip()
        if not part or part == '无':
            continue
        school_name, province = parse_school_and_province(part)
        if school_name:
            results.append((school_name, province))
    return results

# ==================== 3. Provincial Offering University Statistics ====================
def compute_province_schools(df):
    all_provinces = [
        '北京市', '上海市', '天津市', '重庆市',
        '河北省', '山西省', '辽宁省', '吉林省', '黑龙江省',
        '江苏省', '浙江省', '安徽省', '福建省', '江西省', '山东省', '河南省',
        '湖北省', '湖南省', '广东省', '海南省',
        '四川省', '贵州省', '云南省', '陕西省', '甘肃省', '青海省',
        '内蒙古自治区', '广西壮族自治区', '西藏自治区', '宁夏回族自治区', '新疆维吾尔自治区'
    ]
    province_schools = defaultdict(set)
    for _, row in df.iterrows():
        schools = parse_schools_from_row(row['设课大学'])
        for school_name, province in schools:
            if province:
                province_schools[province].add(school_name)
    data = []
    for prov in all_provinces:
        data.append({'省份': prov, '开课高校数': len(province_schools.get(prov, set()))})
    return pd.DataFrame(data), province_schools

# ==================== 4. Spatial Weight Matrix and Moran's I ====================
def generate_spatial_weight():
    provinces = [
        '北京市', '上海市', '天津市', '重庆市',
        '河北省', '山西省', '辽宁省', '吉林省', '黑龙江省',
        '江苏省', '浙江省', '安徽省', '福建省', '江西省', '山东省', '河南省',
        '湖北省', '湖南省', '广东省', '海南省',
        '四川省', '贵州省', '云南省', '陕西省', '甘肃省', '青海省',
        '内蒙古自治区', '广西壮族自治区', '西藏自治区', '宁夏回族自治区', '新疆维吾尔自治区'
    ]
    adjacency = {
        '北京市': ['天津市', '河北省'],
        '上海市': ['江苏省', '浙江省'],
        '天津市': ['北京市', '河北省'],
        '重庆市': ['四川省', '贵州省', '湖北省', '湖南省'],
        '河北省': ['北京市', '天津市', '山西省', '山东省', '河南省', '辽宁省', '内蒙古自治区'],
        '山西省': ['河北省', '陕西省', '河南省', '内蒙古自治区'],
        '辽宁省': ['吉林省', '内蒙古自治区', '河北省'],
        '吉林省': ['辽宁省', '黑龙江省', '内蒙古自治区'],
        '黑龙江省': ['吉林省', '内蒙古自治区'],
        '江苏省': ['上海市', '浙江省', '安徽省', '山东省'],
        '浙江省': ['上海市', '江苏省', '安徽省', '江西省', '福建省'],
        '安徽省': ['江苏省', '浙江省', '江西省', '湖北省', '河南省', '山东省'],
        '福建省': ['浙江省', '江西省', '广东省'],
        '江西省': ['浙江省', '安徽省', '湖北省', '湖南省', '广东省', '福建省'],
        '山东省': ['河北省', '河南省', '江苏省', '安徽省'],
        '河南省': ['河北省', '山西省', '陕西省', '湖北省', '安徽省', '山东省'],
        '湖北省': ['河南省', '安徽省', '江西省', '湖南省', '重庆市', '陕西省'],
        '湖南省': ['湖北省', '江西省', '广东省', '广西壮族自治区', '贵州省', '重庆市'],
        '广东省': ['福建省', '江西省', '湖南省', '广西壮族自治区', '海南省'],
        '海南省': ['广东省'],
        '四川省': ['重庆市', '贵州省', '云南省', '西藏自治区', '青海省', '甘肃省', '陕西省'],
        '贵州省': ['重庆市', '四川省', '云南省', '广西壮族自治区', '湖南省'],
        '云南省': ['四川省', '贵州省', '广西壮族自治区', '西藏自治区'],
        '陕西省': ['山西省', '河南省', '湖北省', '重庆市', '四川省', '甘肃省', '宁夏回族自治区', '内蒙古自治区'],
        '甘肃省': ['陕西省', '四川省', '青海省', '新疆维吾尔自治区', '内蒙古自治区', '宁夏回族自治区'],
        '青海省': ['甘肃省', '四川省', '西藏自治区', '新疆维吾尔自治区'],
        '内蒙古自治区': ['黑龙江省', '吉林省', '辽宁省', '河北省', '山西省', '陕西省', '宁夏回族自治区', '甘肃省'],
        '广西壮族自治区': ['湖南省', '广东省', '贵州省', '云南省'],
        '西藏自治区': ['新疆维吾尔自治区', '青海省', '四川省', '云南省'],
        '宁夏回族自治区': ['陕西省', '甘肃省', '内蒙古自治区'],
        '新疆维吾尔自治区': ['甘肃省', '青海省', '西藏自治区']
    }
    n = len(provinces)
    W = np.zeros((n, n))
    for i, prov in enumerate(provinces):
        for neighbor in adjacency.get(prov, []):
            if neighbor in provinces:
                j = provinces.index(neighbor)
                W[i, j] = 1
    return W, provinces

def calculate_moran_i(x, W):
    n = len(x)
    x_mean = np.mean(x)
    row_sums = W.sum(axis=1, keepdims=True)
    row_sums[row_sums == 0] = 1
    W_norm = W / row_sums
    numerator = 0
    denominator = 0
    for i in range(n):
        for j in range(n):
            numerator += W_norm[i, j] * (x[i] - x_mean) * (x[j] - x_mean)
        denominator += (x[i] - x_mean) ** 2
    if denominator == 0:
        return 0, 0, 0
    I = n / np.sum(W_norm) * numerator / denominator
    E_I = -1 / (n - 1)
    var_I = (n ** 2 * np.sum(W_norm ** 2) - n * np.sum(W_norm ** 2) + 3 * np.sum(W_norm) ** 2) / (
                np.sum(W_norm) ** 2 * (n ** 2 - 1))
    Z = (I - E_I) / np.sqrt(var_I) if var_I > 0 else 0
    p_value = 2 * (1 - stats.norm.cdf(abs(Z)))
    return I, p_value, Z

# ==================== 5. Native-Origin Match Rate Calculation ====================
def parse_province_set(region_str):
    if pd.isna(region_str) or not region_str:
        return set()
    parts = re.split('[、，,;；]', str(region_str))
    provinces = set()
    for p in parts:
        p = p.strip()
        if not p:
            continue
        if len(p) >= 2:
            if p.startswith('内蒙古'):
                provinces.add('内蒙古自治区')
            elif p.startswith('广西'):
                provinces.add('广西壮族自治区')
            elif p.startswith('西藏'):
                provinces.add('西藏自治区')
            elif p.startswith('宁夏'):
                provinces.add('宁夏回族自治区')
            elif p.startswith('新疆'):
                provinces.add('新疆维吾尔自治区')
            elif p[:2] in ['北京', '上海', '天津', '重庆']:
                provinces.add(p[:2] + '市')
            else:
                provinces.add(p[:2] + '省')
    return provinces

def compute_match_rate(df):
    category_map = {
        '少林功夫': '武术流派类', '武当武术': '武术流派类', '回族重刀武术': '武术流派类',
        '沧州武术': '武术流派类', '峨眉武术': '武术流派类', '精武武术': '武术流派类',
        '梅山武术': '武术流派类', '青城武术': '武术流派类', '布依族武术': '武术流派类',
        '围棋': '棋牌类', '象棋': '棋牌类', '蒙古族象棋': '棋牌类', '藏棋': '棋牌类',
        '迎罗汉': '民俗表演类', '迎大旗': '民俗表演类', '泉州刣狮': '民俗表演类',
        '井冈山全堂狮灯': '民俗表演类',
        '华佗五禽戏': '健身气功与器械类', '撂石锁': '健身气功与器械类',
        '十八般武艺': '健身气功与器械类', '沙河藤牌阵': '健身气功与器械类',
        '青萍剑（贾氏青萍剑）': '健身气功与器械类', '易筋经（天台山易筋经）': '健身气功与器械类'
    }
    boxing_keywords = ['拳', '掌', '捶', '肘', '门', '螳螂', '八卦', '形意', '八极', '心意',
                       '五祖', '查', '苌家', '岳家', '蔡李佛', '通背', '戳脚', '咏春', '少北',
                       '西凉', '六合', '巫家', '岩鹰', '莫家', '无极', '徐家', '佛汉', '孙膑',
                       '地术', '通背缠', '拦手', '鹰爪翻子']

    results = []
    for _, row in df.iterrows():
        project_name = row['项目名称']
        region_str = row['申报地区']
        schools_str = row['设课大学']

        if project_name in category_map:
            category = category_map[project_name]
        elif any(k in project_name for k in boxing_keywords):
            category = '拳术类'
        else:
            category = '民俗体育竞技类'

        native_provinces = parse_province_set(region_str)
        if not native_provinces:
            continue

        school_entries = parse_schools_from_row(schools_str)
        total_schools = len(school_entries)
        native_schools = 0
        for _, province in school_entries:
            if province in native_provinces:
                native_schools += 1

        mr = native_schools / total_schools * 100 if total_schools > 0 else 0
        results.append({
            '项目名称': project_name,
            '类别': category,
            '总开课高校数': total_schools,
            '原生地开课高校数': native_schools,
            '原生地匹配率(%)': round(mr, 2)
        })
    return pd.DataFrame(results)

# ==================== 6. University Type Stratified Statistics ====================
def infer_university_type(name):
    if any(k in name for k in ['体育学院', '体育大学', '体育职业']):
        return '体育类'
    if any(k in name for k in ['师范大学', '师范学院', '师范高等专科']):
        return '师范类'
    if any(k in name for k in ['民族大学', '民族学院']):
        return '民族类'
    if any(k in name for k in ['大学', '学院', '职业技术', '职业']):
        return '综合类'
    return '其他'

def compute_university_type_stats(df):
    school_type = {}
    for _, row in df.iterrows():
        school_entries = parse_schools_from_row(row['设课大学'])
        for school_name, _ in school_entries:
            if school_name not in school_type:
                school_type[school_name] = infer_university_type(school_name)

    type_proportions = {
        '综合类': 0.45, '理工类': 0.28, '师范类': 0.14,
        '民族类': 0.03, '体育类': 0.02, '其他': 0.08
    }
    total_schools = 2919
    type_total = {t: int(total_schools * p) for t, p in type_proportions.items()}

    type_open = defaultdict(int)
    for school, typ in school_type.items():
        type_open[typ] += 1
    data = []
    for typ in ['体育类', '民族类', '师范类', '综合类', '理工类', '其他']:
        rate = type_open.get(typ, 0) / type_total.get(typ, 1) * 100
        data.append({
            '高校类型': typ,
            '开课高校数': type_open.get(typ, 0),
            '总高校数': type_total.get(typ, 1),
            '开设率(%)': round(rate, 2)
        })
    return pd.DataFrame(data)

# ==================== 7. Visualization ====================
def plot_province_hbar(df_province, output_file='Figure1_Province_Schools.png'):
    df_sorted = df_province.sort_values('开课高校数', ascending=True)

    fig, ax = plt.subplots(figsize=(12, 10))
    colors = []
    for v in df_sorted['开课高校数']:
        if v >= 10:
            colors.append('#1B3A5C')
        elif v >= 5:
            colors.append('#2C5F8A')
        elif v >= 2:
            colors.append('#6B9BC4')
        else:
            colors.append('#A9C4D9')

    bars = ax.barh(df_sorted['省份'], df_sorted['开课高校数'], color=colors, edgecolor='white', linewidth=0.8, height=0.7)

    for bar, val in zip(bars, df_sorted['开课高校数']):
        if val > 0:
            ax.text(bar.get_width() + 0.3, bar.get_y() + bar.get_height() / 2,
                    f'{int(val)}', va='center', fontsize=9, fontweight='bold')

    ax.set_xlabel('Number of Universities Offering ICH Sports Courses', fontsize=13)
    ax.set_ylabel('Province', fontsize=13)
    ax.set_title('Provincial Distribution of ICH Traditional Sports Courses in Universities', fontsize=15,
                 fontweight='bold')
    ax.grid(axis='x', linestyle=':', alpha=0.4)

    plt.tight_layout()
    plt.savefig(output_file, dpi=300, bbox_inches='tight')
    plt.savefig(output_file.replace('.png', '.pdf'), bbox_inches='tight')
    print(f"✅ Figure 1 saved: {output_file}")

def plot_type_bar(df_type, output_file='Figure2_University_Type_Rate.png'):
    df_sorted = df_type[df_type['高校类型'] != '其他'].sort_values('开设率(%)', ascending=False)

    fig, ax = plt.subplots(figsize=(9, 6))
    colors = ['#1B3A5C', '#2C5F8A', '#4A7B9D', '#7AA5C9', '#A9C4D9']
    bars = ax.bar(df_sorted['高校类型'], df_sorted['开设率(%)'], color=colors, edgecolor='white', linewidth=1)

    for bar, val in zip(bars, df_sorted['开设率(%)']):
        ax.text(bar.get_x() + bar.get_width() / 2, bar.get_height() + 0.8,
                f'{val:.1f}%', ha='center', va='bottom', fontsize=11, fontweight='bold')

    ax.set_xlabel('University Type', fontsize=13)
    ax.set_ylabel('Opening Rate (%)', fontsize=13)
    ax.set_title('ICH Sports Course Opening Rate by University Type', fontsize=15, fontweight='bold')
    ax.set_ylim(0, 75)
    ax.grid(axis='y', linestyle=':', alpha=0.4)

    plt.tight_layout()
    plt.savefig(output_file, dpi=300, bbox_inches='tight')
    plt.savefig(output_file.replace('.png', '.pdf'), bbox_inches='tight')
    print(f"✅ Figure 2 saved: {output_file}")

# ==================== 8. Save All Data to Excel (Revised Version with df_raw Parameter) ====================
def save_all_data(df_raw, df_province, df_mr, df_type, I, p_val, Z, output_file='Q2_All_Data.xlsx'):
    with pd.ExcelWriter(output_file, engine='openpyxl') as writer:
        df_province.to_excel(writer, sheet_name='省份开课高校统计', index=False)
        df_mr.to_excel(writer, sheet_name='原生地匹配率', index=False)
        df_type.to_excel(writer, sheet_name='高校类型开设率', index=False)

        moran_df = pd.DataFrame({
            '统计量': ['Moran\'s I', 'Z-score', 'p-value'],
            '数值': [round(I, 4), round(Z, 4), round(p_val, 4)]
        })
        moran_df.to_excel(writer, sheet_name='Moran检验', index=False)

        mr_summary = df_mr.groupby('类别')['原生地匹配率(%)'].agg(['mean', 'std', 'count']).reset_index()
        mr_summary.columns = ['类别', 'MR均值(%)', 'MR标准差', '项目数']
        mr_summary['MR均值(%)'] = mr_summary['MR均值(%)'].round(2)
        mr_summary['MR标准差'] = mr_summary['MR标准差'].round(2)
        mr_summary.to_excel(writer, sheet_name='各类别MR汇总', index=False)

        t_results = []
        for cat in mr_summary['类别'].unique():
            cat_data = df_mr[df_mr['类别'] == cat]['原生地匹配率(%)'].dropna()
            if len(cat_data) > 1:
                t_stat, p_t = stats.ttest_1samp(cat_data, 3.23)
            else:
                t_stat, p_t = np.nan, np.nan
            t_results.append({
                '类别': cat,
                '样本数': len(cat_data),
                '均值(%)': cat_data.mean().round(2),
                't统计量': round(t_stat, 4) if not np.isnan(t_stat) else np.nan,
                'p值': round(p_t, 4) if not np.isnan(p_t) else np.nan
            })
        df_t = pd.DataFrame(t_results)
        df_t.to_excel(writer, sheet_name='MR_t检验_相对3.23%', index=False)

        # School-Province mapping table (using df_raw)
        school_province_data = []
        for _, row in df_raw.iterrows():
            school_entries = parse_schools_from_row(row['设课大学'])
            for school_name, province in school_entries:
                if school_name and province:
                    school_province_data.append({'项目': row['项目名称'], '学校': school_name, '省份': province})
        if school_province_data:
            df_school_prov = pd.DataFrame(school_province_data)
            df_school_prov = df_school_prov.drop_duplicates()
            df_school_prov.to_excel(writer, sheet_name='学校_省份对照表', index=False)

    print(f"✅ All data saved: {output_file}")

# ==================== 9. Main Program ====================
def main():
    input_file = '全国非遗传统体育项目高校开课情况最终汇总表.xlsx'
    if not os.path.exists(input_file):
        print(f"❌ Error: Input file {input_file} not found")
        return

    print("=" * 60)
    print("Q2: Spatial Differentiation Pattern and Native-Origin Effect Analysis (Based on Precise Annotation)")
    print("=" * 60)

    df = load_data(input_file)
    print(f"✅ Loaded projects: {len(df)}")

    df_province, province_schools = compute_province_schools(df)
    print("✅ Provincial statistics completed")
    print(f"   Total offering universities: {df_province['开课高校数'].sum()}")

    W, provinces = generate_spatial_weight()
    x = np.array([df_province[df_province['省份'] == p]['开课高校数'].values[0]
                  if p in df_province['省份'].values else 0 for p in provinces])
    I, p_val, Z = calculate_moran_i(x, W)
    print(f"📊 Moran's I = {I:.4f}, p = {p_val:.4f}, Z = {Z:.4f}")

    df_mr = compute_match_rate(df)
    mr_mean = df_mr['原生地匹配率(%)'].mean()
    mr_std = df_mr['原生地匹配率(%)'].std()
    print("✅ Native-origin match rate calculation completed")
    print(f"   MR Mean: {mr_mean:.2f}%, SD: {mr_std:.2f}%")

    t_stat, p_t = stats.ttest_1samp(df_mr['原生地匹配率(%)'], 3.23)
    print(f"   One-sample t-test (vs 3.23%): t={t_stat:.4f}, p={p_t:.4f}")

    df_type = compute_university_type_stats(df)
    print("✅ University type statistics completed")

    # Pass df as df_raw
    save_all_data(df, df_province, df_mr, df_type, I, p_val, Z)

    plot_province_hbar(df_province)
    plot_type_bar(df_type)

    print("\n" + "=" * 60)
    print("🎉 All tasks completed! Generated files:")
    print("   - Q2_All_Data.xlsx (7 sheets included)")
    print("   - Figure1_Province_Schools.png/pdf")
    print("   - Figure2_University_Type_Rate.png/pdf")
    print("=" * 60)

if __name__ == '__main__':
    main()
